Skip to main content

3.4 — Multi Neuron Training

5.5 Training Multiple Neurons​

A layer with multiple neurons has multiple weights and biases.

y=Wx+by=Wx+b

For 2 neurons and 3 inputs:

W=[w11w12w13w21w22w23]W= \begin{bmatrix} w_{11}&w_{12}&w_{13}\\ w_{21}&w_{22}&w_{23} \end{bmatrix}

Each weight can be updated independently.


5.7 Updating the Weights​

Each neuron has its own gradients. calculated by d (Loss)/d (Weight)

NeuronNeuron 1 GradientNeuron 2 Gradient
weight 1dLdw11=−2x1(error1)\frac{dL}{dw_{11}}=-2x_1(error_1)dLdw21=−2x1(error2)\frac{dL}{dw_{21}}=-2x_1(error_2)
weight 2dLdw12=−2x2(error1)\frac{dL}{dw_{12}}=-2x_2(error_1)dLdw22=−2x2(error2)\frac{dL}{dw_{22}}=-2x_2(error_2)
The same gradient-descent rule applies:
new.parameter=old.parameter−η×gradient\boxed{new.parameter=old.parameter-\eta\times gradient}

So multiple neurons simply mean more parameters to update.


5.8 Python Example​

Check [Matrix Outer Product](../../Course-1-Mathematics and Frameworks/Ch-1 Maths -Linear-Algebra/4 Matrix Outer Product.mdx) for more info on the line np.outer(error, x).
Check [Matrix Multiplication](../../Course-1-Mathematics and Frameworks/Ch-1 Maths -Linear-Algebra/5 Matrix-Dot Product.mdx) for more info on the line W @ x.

import numpy as np

x = np.array([2, 3])
target = np.array([10, 15])

W = np.array([
[1.0, 1.0],
[1.0, 1.0]
])

b = np.array([0.0, 0.0])

learning_rate = 0.01
printed=False
for step in range(100):
prediction = W @ x + b
error = target - prediction
loss = np.sum(error ** 2)

gradient_W = -2 * np.outer(error, x)
if printed==False:
print("error is",error," | inputs are",x," | gradient is",gradient_W)
printed = True
#this means we are first putting x1 as value and calculating both errors and then we are putting x2 as value and calculating both errors

#so technically we cross multiplied every error to every neuron. eg, 2,3 * error on n1 means when n1 gets inputs as 2 and 3, it has error = error on n1
gradient_b = -2 * error

W = W - learning_rate * gradient_W
b = b - learning_rate * gradient_b

print("prediction:", prediction,"loss:", loss,"W:", W,"b:", b)